DeConv2DGradFilter
计算反卷积(转置卷积)对权重的梯度。dy、x 为 NHWC。
实现上先对 dy 做 im2row(image to row,图像块按行展开成矩阵),再与 x 做 GEMM(GEneral Matrix Multiply,通用矩阵乘),并按 batch、group 累加得到 dw。
令 \(C_o' = C_{out}/G\),\(C_i' = C_{in}/G\)。对输入空间位置 \((i_h, i_w)\)、核偏移 \((k_h, k_w)\),对应 dy 坐标为:
\[\begin{split}\begin{aligned}
o_h &= i_h \cdot s_h + k_h \cdot d_h - pad_u \\
o_w &= i_w \cdot s_w + k_w \cdot d_w - pad_l
\end{aligned}\end{split}\]
当 \((o_h, o_w)\) 落在 dy 有效范围内时取值,否则为 0。
dw 按组展平,元素总数为 \(C_{out} \cdot K_h \cdot K_w \cdot C_i'\)。
- 输入:
dy_data - 输出梯度地址,形状
[batch, out_h, out_w, out_c]x_data - 前向输入地址,形状
[batch, in_h, in_w, in_c]param -
long long参数数组地址,字段见下core_mask - 核掩码(仅共享存储版本)
param 布局(``long long param[20]``):
1/* 每个槽位 8 字节;param[17] 存工作区指针 */
2param[0] = input_batch;
3param[1] = in_h;
4param[2] = in_w;
5param[3] = in_c;
6param[4] = output_batch; /* 算子循环使用该值;通常与 input_batch 相同 */
7param[5] = out_h;
8param[6] = out_w;
9param[7] = out_c;
10param[8] = kernel_h;
11param[9] = kernel_w;
12param[10] = stride_h;
13param[11] = stride_w;
14param[12] = pad_u;
15param[13] = pad_l;
16param[14] = dilation_h;
17param[15] = dilation_w;
18param[16] = group;
19param[17] = (long long)workspace;
20/* 工作区指针,元素数见下;再乘 sizeof(dtype) 为字节数。
21 * 令 mat = (out_c/G)*Kh*Kw*in_h*in_w,dw = group*(out_c/G)*Kh*Kw*(in_c/G):
22 * 私有存储 (_p):mat(仅 im2row)
23 * 共享存储 (_s):core_num*mat + (core_num-1)*dw
24 * 前段每核一份 im2row;后段给核 1..N-1 各一份局部 dw,规约进 dw_data
25 */
- 输出:
dw_data - 权重梯度;算子内部会先清零再累加
- 支持平台:
FT78NEMT7004
备注
FT78NE 支持 fp32
MT7004 支持 fp16、fp32
workspace 元素数见
param[17]注释;按元素数乘sizeof(dtype)分配
共享存储版本:
-
void hp_deconv_grad_filter_s(float16 *dy_data, float16 *x_data, float16 *dw_data, long long *param, int core_mask)
-
void fp_deconv_grad_filter_s(float *dy_data, float *x_data, float *dw_data, long long *param, int core_mask)
C调用示例:
1// MT7004 示例(共享存储多核,DDR 地址)
2void TestDeconvGradFilterSMCFp32(int core_mask) {
3 int core_id = get_core_id();
4 int logic_core_id = GetLogicCoreId(core_mask, core_id);
5 int core_num = GetCoreNum(core_mask);
6 float *x = (float *)0x81000000;
7 float *dy = (float *)0x82000000;
8 float *dw = (float *)0x83000000;
9 float *temp_space = (float *)0x85000000;
10 long long params[20];
11 if (logic_core_id == 0) {
12 params[0] = 4;
13 params[1] = 4;
14 params[2] = 4;
15 params[3] = 4;
16 params[4] = 4;
17 params[5] = 4;
18 params[6] = 4;
19 params[7] = 4;
20 params[8] = 4;
21 params[9] = 4;
22 params[10] = 1;
23 params[11] = 1;
24 params[12] = 0;
25 params[13] = 0;
26 params[14] = 1;
27 params[15] = 1;
28 params[16] = 1;
29 params[17] = (long long)temp_space;
30 }
31 sys_bar(0, core_num);
32 fp_deconv_grad_filter_s(dy, x, dw, params, core_mask);
33}
34
35void main() {
36 int core_mask = 0b1111;
37 TestDeconvGradFilterSMCFp32(core_mask);
38}
私有存储版本:
-
void hp_deconv_grad_filter_p(float16 *dy_data, float16 *x_data, float16 *dw_data, long long *param)
-
void fp_deconv_grad_filter_p(float *dy_data, float *x_data, float *dw_data, long long *param)
C调用示例:
1// MT7004 示例(私有存储单核,AM 地址)
2void TestDeconvGradFilterAMFp32(void) {
3 float *x = (float *)0x10010000;
4 float *dy = (float *)0x10020000;
5 float *dw = (float *)0x10030000;
6 float *temp_space = (float *)0x10050000;
7 long long params[20];
8 params[0] = 4;
9 params[1] = 4;
10 params[2] = 4;
11 params[3] = 4;
12 params[4] = 4;
13 params[5] = 4;
14 params[6] = 4;
15 params[7] = 4;
16 params[8] = 4;
17 params[9] = 4;
18 params[10] = 1;
19 params[11] = 1;
20 params[12] = 0;
21 params[13] = 0;
22 params[14] = 1;
23 params[15] = 1;
24 params[16] = 1;
25 params[17] = (long long)temp_space;
26 fp_deconv_grad_filter_p(dy, x, dw, params);
27}
28
29void main() {
30 TestDeconvGradFilterAMFp32();
31}